# ==========================================================
# Table 3 replication (nonlinearity) + Excel export
# 
# FNL (XGBoost): two-stage MI  (Y = a + b*Yhat + eps)
# OLS:           one-stage MI  (Y = b0 + b1*X + b2*Z + eps)
#                directly on covariates, with full posterior
#                draws of coefficients and residuals.
# ==========================================================
install.packages(c("xgboost","writexl","MASS"))  # uncomment if needed
library(xgboost)
library(writexl)
library(MASS)
rm(list = ls())

# Resolve paths without config.R
SCRIPT_DIR <- tryCatch({
  this_file <- normalizePath(sys.frames()[[1]]$ofile)
  dirname(this_file)
}, error = function(e) {
  tryCatch(dirname(rstudioapi::getSourceEditorContext()$path), error = function(e2) getwd())
})

TABLE3_OUT <- normalizePath(file.path(SCRIPT_DIR, "..", "Output"), mustWork = FALSE)



set.seed(42)
N <- 10000
lambdas <- c(0.001, 0.01, 0.1, 0.5)
M_imp <- 10  # MI draws

# -------- Nonlinear link h(s; lambda)
h_bc <- function(s, lambda) {
  if (abs(lambda) < 1e-8) {
    s
  } else {
    (exp(lambda * s) - 1) / lambda
  }
}

# -------- Period generator: Period 1 muZ=0, Period 2 muZ=1
gen_period_bc <- function(N, lambda, muZ) {
  X  <- rnorm(N, 0, 1)
  Z  <- rnorm(N, muZ, 1)
  eU <- rnorm(N, 0, 1)
  U  <- 1 + 2 * Z + eU                 # U_t = 1 + 2 Z_t + eU
  s  <- X + U
  eY <- rnorm(N, 0, 1)
  Y  <- h_bc(s, lambda) + eY           # Y_t = h(X+U) + eY
  data.frame(X, Z, U, Y)
}

# -------- Two-stage MI: calibration Y = a + b * Yhat + eps  (for FNL/XGBoost)
mi_calibrate_and_compute <- function(Y1, Yhat1, Yhat2, tau, M = 10) {
  df  <- length(Y1) - 2
  X1  <- cbind(1, Yhat1)
  fit <- lm.fit(x = X1, y = Y1)
  beta_hat   <- coef(fit)
  sigma2_hat <- sum(residuals(fit)^2) / df
  V_inv <- crossprod(X1)
  V_beta <- function(s2) s2 * solve(V_inv)
  
  X2 <- cbind(1, Yhat2)
  
  means  <- numeric(M)
  W_mean <- numeric(M)
  props  <- numeric(M)
  W_prop <- numeric(M)
  
  for (m in 1:M) {
    sigma2_m <- df * sigma2_hat / rchisq(1, df)
    beta_m   <- MASS::mvrnorm(1, mu = beta_hat, Sigma = V_beta(sigma2_m))
    y_imp    <- as.numeric(X2 %*% beta_m) + rnorm(length(Yhat2), 0, sqrt(sigma2_m))
    
    means[m]  <- mean(y_imp)
    W_mean[m] <- var(y_imp) / length(y_imp)
    
    ind <- (y_imp < tau)
    p_m <- mean(ind)
    props[m]  <- p_m
    W_prop[m] <- p_m * (1 - p_m) / length(ind)
  }
  
  pool <- function(qm, Wm) {
    Qbar <- mean(qm); Wbar <- mean(Wm); B <- var(qm)
    Tvar <- Wbar + (1 + 1/M) * B
    list(est = Qbar, se = sqrt(Tvar))
  }
  
  out_mean <- pool(means, W_mean)
  out_prop <- pool(props, W_prop)
  
  list(
    mean_est   = out_mean$est,
    mean_se    = out_mean$se,
    pov_est_pp = 100 * out_prop$est,
    pov_se_pp  = 100 * out_prop$se
  )
}

# -------- One-stage MI: direct OLS  Y = b0 + b1*X + b2*Z + eps
#          Draws sigma2 from scaled inv-chi2, beta from MVN,
#          then imputes Y for Period 2 using Period-2 covariates.
mi_direct_ols <- function(dat_train, dat_predict, tau, M = 10) {
  # Fit OLS on training data (Period 1)
  fit <- lm(Y ~ X + Z, data = dat_train)
  beta_hat   <- coef(fit)
  sigma2_hat <- summary(fit)$sigma^2
  df         <- fit$df.residual
  V_beta_hat <- vcov(fit)   # = sigma2_hat * (X'X)^{-1}
  # We need V_beta as a function of sigma2:
  # vcov(fit) = sigma2_hat * (X'X)^{-1}
  # so (X'X)^{-1} = vcov(fit) / sigma2_hat
  XtXinv <- V_beta_hat / sigma2_hat
  
  # Design matrix for Period 2
  X2 <- model.matrix(~ X + Z, data = dat_predict)
  N2 <- nrow(dat_predict)
  
  means  <- numeric(M)
  W_mean <- numeric(M)
  props  <- numeric(M)
  W_prop <- numeric(M)
  
  for (m in 1:M) {
    # Draw sigma2 from its posterior: scaled inverse-chi-squared
    sigma2_m <- df * sigma2_hat / rchisq(1, df)
    
    # Draw beta from its posterior: MVN(beta_hat, sigma2_m * (X'X)^{-1})
    beta_m <- MASS::mvrnorm(1, mu = beta_hat, Sigma = sigma2_m * XtXinv)
    
    # Impute Y for Period 2
    mu_m  <- as.numeric(X2 %*% beta_m)
    y_imp <- mu_m + rnorm(N2, 0, sqrt(sigma2_m))
    
    # Mean of imputed Y
    means[m]  <- mean(y_imp)
    W_mean[m] <- var(y_imp) / N2
    
    # Poverty indicator
    ind <- (y_imp < tau)
    p_m <- mean(ind)
    props[m]  <- p_m
    W_prop[m] <- p_m * (1 - p_m) / N2
  }
  
  pool <- function(qm, Wm) {
    Qbar <- mean(qm); Wbar <- mean(Wm); B <- var(qm)
    Tvar <- Wbar + (1 + 1/M) * B
    list(est = Qbar, se = sqrt(Tvar))
  }
  
  out_mean <- pool(means, W_mean)
  out_prop <- pool(props, W_prop)
  
  list(
    mean_est   = out_mean$est,
    mean_se    = out_mean$se,
    pov_est_pp = 100 * out_prop$est,
    pov_se_pp  = 100 * out_prop$se
  )
}

# -------- XGBoost fitting helper
fit_xgb <- function(Xmat, y) {
  dtrain <- xgb.DMatrix(data = Xmat, label = y)
  params <- list(
    objective = "reg:squarederror",
    eta = 0.1, max_depth = 4,
    subsample = 0.8, colsample_bytree = 0.8,
    min_child_weight = 1,
    seed = 42
  )
  xgb.train(params = params, data = dtrain, nrounds = 300, verbose = 0)
}

rows <- list()

for (i in seq_along(lambdas)) {
  lambda <- lambdas[i]
  # Reset seed per iteration so XGBoost version differences don't

  # cascade through the RNG state to subsequent lambdas.
  set.seed(42 + i)

  # --- simulate periods
  dat1 <- gen_period_bc(N, lambda, muZ = 0)
  dat2 <- gen_period_bc(N, lambda, muZ = 1)
  
  # quick diagnostics
  cat(sprintf("lambda=%.3f | mean(Y1)=%.2f, mean(Y2)=%.2f, mean(Z2)=%.2f\n",
              lambda, mean(dat1$Y), mean(dat2$Y), mean(dat2$Z)))
  
  # poverty line = median(Y1)
  tau <- median(dat1$Y)
  
  # TRUE period-2 stats
  true_mean   <- mean(dat2$Y)
  true_se     <- sd(dat2$Y) / sqrt(N)
  p_true      <- mean(dat2$Y < tau)
  se_true_pp  <- 100 * sqrt(p_true * (1 - p_true) / N)
  
  add_row <- function(model, proxy, meanY, seY, pov_pp, se_pp) {
    data.frame(
      Exponent     = lambda,
      Model        = model,
      Proxy        = proxy,
      `Mean Y`     = meanY,
      `SE (mean)`  = seY,
      `Poverty (%)`= pov_pp,
      `SE (pp)`    = se_pp,
      check.names  = FALSE
    )
  }
  
  # FNL (xgboost) WITHOUT Z  -- two-stage calibration MI
  X1_noZ <- as.matrix(dat1[, "X", drop = FALSE])
  X2_noZ <- as.matrix(dat2[, "X", drop = FALSE])
  xgb_noZ <- fit_xgb(X1_noZ, dat1$Y)
  yhat1_noZ <- as.numeric(predict(xgb_noZ, xgb.DMatrix(X1_noZ)))
  yhat2_noZ <- as.numeric(predict(xgb_noZ, xgb.DMatrix(X2_noZ)))
  set.seed(100 * i + 1)  # isolate MI draws from XGBoost RNG consumption
  mi_noZ <- mi_calibrate_and_compute(dat1$Y, yhat1_noZ, yhat2_noZ, tau, M = M_imp)

  # FNL (xgboost) WITH Z  -- two-stage calibration MI
  X1_wZ <- as.matrix(dat1[, c("X","Z")])
  X2_wZ <- as.matrix(dat2[, c("X","Z")])
  xgb_wZ <- fit_xgb(X1_wZ, dat1$Y)
  yhat1_wZ <- as.numeric(predict(xgb_wZ, xgb.DMatrix(X1_wZ)))
  yhat2_wZ <- as.numeric(predict(xgb_wZ, xgb.DMatrix(X2_wZ)))
  set.seed(100 * i + 2)  # isolate MI draws from XGBoost RNG consumption
  mi_wZ <- mi_calibrate_and_compute(dat1$Y, yhat1_wZ, yhat2_wZ, tau, M = M_imp)

  # OLS WITH Z  -- one-stage direct MI on Y ~ X + Z
  set.seed(100 * i + 3)  # isolate MI draws from XGBoost RNG consumption
  mi_ols <- mi_direct_ols(dat_train = dat1, dat_predict = dat2, tau = tau, M = M_imp)
  
  # collect rows
  rows[[length(rows)+1]] <- add_row("TRUE", "N/A",
                                    meanY = true_mean, seY = true_se,
                                    pov_pp = 100 * p_true, se_pp = se_true_pp)
  
  rows[[length(rows)+1]] <- add_row("FNL", "No",
                                    meanY = mi_noZ$mean_est, seY = mi_noZ$mean_se,
                                    pov_pp = mi_noZ$pov_est_pp, se_pp = mi_noZ$pov_se_pp)
  
  rows[[length(rows)+1]] <- add_row("FNL", "Yes",
                                    meanY = mi_wZ$mean_est, seY = mi_wZ$mean_se,
                                    pov_pp = mi_wZ$pov_est_pp, se_pp = mi_wZ$pov_se_pp)
  
  rows[[length(rows)+1]] <- add_row("OLS", "Yes",
                                    meanY = mi_ols$mean_est, seY = mi_ols$mean_se,
                                    pov_pp = mi_ols$pov_est_pp, se_pp = mi_ols$pov_se_pp)
}

tab3 <- do.call(rbind, rows)

# optional rounding
tab3$`Mean Y`     <- round(tab3$`Mean Y`, 2)
tab3$`SE (mean)`  <- round(tab3$`SE (mean)`, 2)
tab3$`Poverty (%)`<- round(tab3$`Poverty (%)`, 1)
tab3$`SE (pp)`    <- round(tab3$`SE (pp)`, 2)

# write to Excel
dir.create(TABLE3_OUT, recursive = TRUE, showWarnings = FALSE)
write_xlsx(list("Table 3" = tab3), path = file.path(TABLE3_OUT, "Table3_results.xlsx"))
cat("Wrote:", file.path(TABLE3_OUT, "Table3_results.xlsx"), "\n")
